DivGrad
计算逐元素除法 \(Y = x_1 / x_2\) 的反向梯度。支持同形逐元素,以及 x1 / x2 广播(一侧尺寸更大)的场景。
同形时:
\[\begin{split}\begin{aligned}
dx_1 &= \frac{dy}{x_2} \\
dx_2 &= -\frac{dy \cdot x_1}{x_2^{2}}
\end{aligned}\end{split}\]
广播时,先在广播后的公共形状上按上式计算,再对较小输入一侧沿广播轴做 ReduceSum,使 dx1 / dx2 与对应输入同形。
算子内部按 x1_size 与 x2_size 比较自动分发:x1 更大、x2 更大或两者相等。
- 输入:
dy - 上游梯度地址
x1_data - 前向被除数地址
x2_data - 前向除数地址
params -
Parameter结构体地址,布局见下core_mask - 核掩码(仅共享存储版本使用)
Parameter 定义:
1typedef struct {
2 void *tile_data0; // 临时工作空间 0,长度至少 dy_size 个元素
3 void *tile_data1; // 临时工作空间 1,长度至少 dy_size 个元素
4 void *tile_data2; // 临时工作空间 2,长度至少 dy_size 个元素
5 int *large_shape; // 较大输入形状,长度 ndims
6 int *small_shape; // 较小输入形状,长度 ndims
7 int *out_shape; // 输出 / dy 形状,长度 ndims
8 long long ndims; // 维数
9 long long dy_size; // dy 元素个数
10 long long x1_size; // x1 元素个数
11 long long x2_size; // x2 元素个数
12 int *large_strides; // 较大输入 strides,长度 ndims
13 int *small_strides; // 较小输入 strides,长度 ndims
14 int *out_strides; // 输出 strides,长度 ndims
15 int *large_multiples; // 较大输入广播倍数,长度 ndims
16 int *small_multiples; // 较小输入广播倍数,长度 ndims
17 int *indices; // 广播临时索引,长度 ndims,须清零
18 int *x1_shape; // x1 形状,长度 ndims
19 int *x2_shape; // x2 形状,长度 ndims
20} Parameter;
- 输出:
dx1 - 对
x1的梯度,形状与x1相同dx2 - 对
x2的梯度,形状与x2相同
- 支持平台:
FT78NEMT7004
备注
FT78NE 支持 fp32
MT7004 支持 fp16、fp32
tile_data0/tile_data1/tile_data2各需至少dy_size个元素;fp32 为dy_size * sizeof(float)字节,fp16 为dy_size * sizeof(float16)字节*_strides/*_multiples由算子内部使用,调用方分配长度ndims的int数组即可;indices长度ndims,须预先清零large_shape/small_shape/out_shape按x1_size、x2_size关系设置:较大侧对应large_shape,out_shape取较大形状
共享存储版本:
-
void hp_div_grad_s(float16 *dy, float16 *dx1, float16 *dx2, float16 *x1_data, float16 *x2_data, Parameter *params, int core_mask)
-
void fp_div_grad_s(float *dy, float *dx1, float *dx2, float *x1_data, float *x2_data, Parameter *params, int core_mask)
C调用示例:
1// MT7004 示例(共享存储多核,DDR 地址)
2void TestDivGradSMCFp32(int core_mask) {
3 int core_id = get_core_id();
4 int logic_core_id = GetLogicCoreId(core_mask, core_id);
5 int core_num = GetCoreNum(core_mask);
6 float *dy = (float *)0x81000000;
7 float *dx1 = (float *)0x82000000;
8 float *dx2 = (float *)0x83000000;
9 float *x1_data = (float *)0x84000000;
10 float *x2_data = (float *)0x85000000;
11 float *tile0 = (float *)0x86000000;
12 float *tile1 = (float *)0x87000000;
13 float *tile2 = (float *)0x88000000;
14 int *strides_base = (int *)0x8B000000;
15 int x1_shape[4] = {4, 4, 4, 4};
16 int x2_shape[4] = {4, 4, 4, 4};
17 Parameter params;
18 if (logic_core_id == 0) {
19 params.tile_data0 = tile0;
20 params.tile_data1 = tile1;
21 params.tile_data2 = tile2;
22 params.large_shape = x1_shape;
23 params.small_shape = x2_shape;
24 params.out_shape = x1_shape;
25 params.ndims = 4;
26 params.dy_size = 256;
27 params.x1_size = 256;
28 params.x2_size = 256;
29 params.large_strides = strides_base;
30 params.small_strides = strides_base + 4;
31 params.out_strides = strides_base + 8;
32 params.large_multiples = strides_base + 12;
33 params.small_multiples = strides_base + 16;
34 params.indices = strides_base + 20;
35 params.x1_shape = x1_shape;
36 params.x2_shape = x2_shape;
37 memset(params.indices, 0, 4 * sizeof(int));
38 }
39 sys_bar(0, core_num);
40 fp_div_grad_s(dy, dx1, dx2, x1_data, x2_data, ¶ms, core_mask);
41}
42
43void main() {
44 int core_mask = 0b1111;
45 TestDivGradSMCFp32(core_mask);
46}
私有存储版本:
-
void hp_div_grad_p(float16 *dy, float16 *dx1, float16 *dx2, float16 *x1_data, float16 *x2_data, Parameter *params)
-
void fp_div_grad_p(float *dy, float *dx1, float *dx2, float *x1_data, float *x2_data, Parameter *params)
C调用示例:
1// MT7004 示例(私有存储单核,AM 地址)
2void TestDivGradAMFp32(void) {
3 float *dy = (float *)0x10010000;
4 float *dx1 = (float *)0x10016000;
5 float *dx2 = (float *)0x10020000;
6 float *x1_data = (float *)0x10026000;
7 float *x2_data = (float *)0x10030000;
8 float *tile0 = (float *)0x10036000;
9 float *tile1 = (float *)0x10040000;
10 float *tile2 = (float *)0x10046000;
11 int *strides_base = (int *)0x10050000;
12 int x1_shape[4] = {4, 4, 4, 4};
13 int x2_shape[4] = {4, 4, 4, 4};
14 Parameter params;
15 params.tile_data0 = tile0;
16 params.tile_data1 = tile1;
17 params.tile_data2 = tile2;
18 params.large_shape = x1_shape;
19 params.small_shape = x2_shape;
20 params.out_shape = x1_shape;
21 params.ndims = 4;
22 params.dy_size = 256;
23 params.x1_size = 256;
24 params.x2_size = 256;
25 params.large_strides = strides_base;
26 params.small_strides = strides_base + 4;
27 params.out_strides = strides_base + 8;
28 params.large_multiples = strides_base + 12;
29 params.small_multiples = strides_base + 16;
30 params.indices = strides_base + 20;
31 params.x1_shape = x1_shape;
32 params.x2_shape = x2_shape;
33 memset(params.indices, 0, 4 * sizeof(int));
34 fp_div_grad_p(dy, dx1, dx2, x1_data, x2_data, ¶ms);
35}
36
37void main() {
38 TestDivGradAMFp32();
39}